iT邦幫忙

2026 iThome 鐵人賽

DAY 13
0
Software Development

GPU 效能優化實戰:30 天從 Kernel 到 Profiling (重賽版)系列 第 13

Day 13|速度沒有贏,那精度可以贏嗎?(重賽版)

  • 分享至 

  • xImage
  •  

2.5條梧曾經說過
https://megapx-assets.dcard.tw/images/c87c5710-213d-42e1-8c48-40695b950abc/orig.jpeg

前幾天我們一直在想一件事:

能不能用 INT8 Tensor Core 做出比 FP64 更快的高精度 GEMM?
原生的硬體贏了
但 額 人家也知道這個方法 NV也不是吃白飯的
速度上 人家的int8 mma利用率高多了 呵呵呵

所以呢

如果目標只是「輸入 double、輸出 double,然後速度贏 cuBLAS」,這件事沒有想像中簡單。cuBLAS 本來就是 NVIDIA 自己長年優化的 GEMM library,尤其現在連 emulation 路線也做得很成熟。

所以

今天來寫 DoubleDouble

昨天我們其實很努力想把它做快。

但 profile 看完以後,現實也很明顯:要直接靠現在這條路去贏 cuBLAS,很難。不是說 INT8 GEMM 不能跑,而是整個流程裡面還有太多 GEMM 前後的東西。你把 MMA kernel 再磨一點,可能有幫助,但它不是現在最大的洞。

所以今天換方向。

速度先不贏,那可不可以把精度往上移?

今天就來寫 DoubleDouble。


DoubleDouble 這次到底要做到什麼

DoubleDouble 的格式很簡單:

struct DoubleDouble {
    double hi;
    double lo;
};

數值語意是:

value = hi + lo

一個 double 大約有 53-bit significand。DoubleDouble 用兩個不重疊的 double 去放一個數,所以有效精度大約是:

53 + 53 = 106 bits

這裡不是要說它是某個硬體原生格式。它不是。

這次要做的事情也不是「宣告一個 struct 就結束」。真正要確認的是:

DoubleDouble 的 lo 有沒有真的進入 residue
residue GEMM 後有沒有被保留下來
CRT reconstruction 後有沒有回到 DoubleDouble output

只要中間任何一步把 hi + lo 直接變回普通 double,這件事就失敗了。


第一個實作點:DoubleDouble 先轉成 scaled integer words

這次最關鍵的地方其實在 GEMM 前面。

不能這樣寫:

double x = dd.hi + dd.lo;

這行會把很多 lo 直接吃掉。

所以我們現在做的是把 DoubleDouble 轉成整數表示。

概念上是:

DoubleDouble(hi, lo)
-> 乘上一個 2^s scale
-> 拆成整數 words

目前單一元素用:

uint64_t words[3];

也就是 192-bit 的暫存整數表示。這裡只是單一 input element 的 scaled integer,不是最後 CRT reconstruction 的總空間。最後重建會用更寬的 UInt384

轉換的核心是對 hilo 分別做:

frexp(x) -> mantissa / exponent
shift = exponent - 53 + scale
mantissa 放進 uint64_t words

正負號也要一起處理。hi 是負數時,整個 scaled integer 要用負的 residue 表示;lo 不能被當成獨立的普通 double 結果丟掉。

這段的重點不是浮點加法,而是:

先把 DoubleDouble 的有效位轉成一個整數,再進入 modulo 世界。

這樣後面每個 modulus 看到的才是同一個高精度整數。


第二個實作點:同一份 words 對所有 modulus 取 residue

一開始最直覺的寫法會變成:

for each modulus:
    DoubleDouble -> integer
    integer -> residue

這樣會讓同一個元素重複做很多次 frexp、round、shift。

現在改成:

for each element:
    DoubleDouble -> words
    for each modulus:
        words -> residue

也就是每個元素只做一次 DoubleDouble 拆解。

words -> residue 的做法是把 64-bit word 再拆成 32-bit chunk,從高位往低位掃:

rem = (rem * 2^32 + chunk) mod p

最後把 residue 轉成 balanced int8:

if rem > p / 2:
    rem -= p

所以 residue matrix 裡面放的不是原始數字,而是每個 modulus 下的代表值。

這裡也解釋了為什麼我們可以用接近 8-bit 的 modulus。像 p=256 時,255 可以用 -1 表示,因為在 mod 256 裡:

-1 == 255

Tensor Core 吃的是 signed int8,但數學上我們用 balanced residue 解釋它。


第三個實作點:modulus count 必須擴上去

DoubleDouble 不是 53-bit target。

它大約是 106-bit target。

GEMM 裡每個輸出元素是 K 個乘積累加,所以 CRT 空間要估:

required_crt_bits = 2 * input_bits + ceil(log2(K))

對 DoubleDouble:

input_bits = 106

所以:

case required bits selected modulus count CRT bits
K=128 219 29 約 220.1
K=4096 224 30 約 226.9

這也是為什麼 20 個 modulus 不夠。

20 個 modulus 大約只有:

155 bits

對 DoubleDouble 的乘加空間來說太小。

目前表裡可用的 modulus 到 49 個,整體 bit budget 大約 341.9 bits。DoubleDouble 不需要全拿;K=4096 目前約 30 個就夠。但重點是這個表要能往上選,而不是停在舊的 20 個。

如果使用者硬指定:

DoubleDouble
K = 4096
num_moduli = 20

現在正確行為是直接拒絕:

Insufficient CRT bits for DoubleDouble

這種拒絕很重要。高精度程式不能「算完再說」,因為 CRT 空間不夠時,它不是誤差大一點而已,而是會 wrap 到另一個整數。


第四個實作點:每個 modulus 做一個 INT8 GEMM

產生 residue matrix 以後,資料長這樣:

A_mod[p][m, k]
B_mod[p][k, n]

對每個 modulus 做:

C_mod[p] = A_mod[p] * B_mod[p]

A_modB_mod 是 int8,accumulator 是 int32。

這裡的 kernel contract 很單純:

input:  int8 residue matrix
output: int32 residue GEMM result

高精度不在單一 GEMM 裡完成。

高精度是靠:

多個 modulus 的 int32 GEMM result
再一起做 CRT reconstruction

所以這次我沒有把文章重點放在 MMA 指令。昨天已經看過 MMA 和 profile 了,今天真正新增的是 DoubleDouble 怎麼進入這條 residue pipeline。


第五個實作點:CRT reconstruction 換成 UInt384

C_mod[p] 算完以後,要把多個 modulus 的結果拼回同一個整數。

DoubleDouble K=4096 需要大約 224-bit CRT 空間,所以舊的 192-bit reconstruction 不夠。

這次用:

UInt384

重建時先算:

M_total = product(p_i)
M_i     = M_total / p_i
inv_i   = inverse(M_i mod p_i)

每個輸出元素再做:

result = sum C_mod[i] * inv_i * M_i
result = result mod M_total

如果:

result > M_total / 2

就把它視為負數區間。

這裡有一個實作細節:result 是很多個 M_i * term 加起來,理論上要做 mod M_total。目前因為項數有限,實作上用重複 subtract M_total 的方式把它壓回範圍內。這不是最快,但對這個 correctness path 夠直接,也比較好 debug。


第六個實作點:UInt384 回到 DoubleDouble

重建完得到的是一個寬整數,還要乘回 scale。

如果這裡只轉成普通 double,前面保住的低位又會不見。

所以現在的做法是把 UInt384 拆成 32-bit chunk,逐段變成 double term,再用 DoubleDouble addition 累加:

for each 32-bit chunk:
    term = chunk * 2^(chunk_offset) * scale
    dd_acc += term

最後得到:

C.hi
C.lo

這段很重要,因為它讓 CRT reconstruction 的結果不是只落回 53-bit,而是回到 DoubleDouble 表示。


驗證:不要只看 random matrix

這種東西不能只拿 random matrix 測。

random input 很容易讓 FP64 看起來也還可以,或者讓錯誤被平均掉。

這次最有用的測試是故意做 low-bit cancellation。

設定:

A = +/-1 + 2^-60
B = 1
K = 128

FP64 collapse 後只看到 +1-1,加起來是:

0

但 DoubleDouble 應該留下每一項的 2^-60

128 * 2^-60 = 1.110223e-16

實測:

FP64 collapsed: 0
DoubleDouble:   1.110223e-16
oracle:         1.110223e-16

這個 case 很直接。它不是在看一般誤差漂不漂亮,而是在問:

FP64 看不到的那一段,我們是不是真的保存下來了?

答案是有。


128³ 的結果

目前 128x128x128 測了三組:

case DoubleDouble max_rel FP64 collapsed max_rel
cancellation_lo_only 0 1.0
constant_lo_1e-20 約 1.64e-33 約 2.00e-20
varied_hi_lo 約 4.04e-27 約 7.92e-12

這張表的重點是對照組。

同一份輸入,如果先 collapse 成 FP64,低位會掉;走 DoubleDouble residue path,低位可以留下。

所以這次可以比較有把握地說:

DoubleDouble correctness path 是成立的

但不能說:

高速 DoubleDouble GEMM 已經完成

這兩句要分開。


profile 看到什麼

這次有把 stage timing 印出來。

128³ 測試裡,INT8 GEMM 本身大概只有:

0.3 - 0.4 ms

但整個 DoubleDouble engine 大約:

23 - 31 ms

主要時間在:

DoubleDouble serialization: 約 13 - 15 ms
UInt384 CRT reconstruction: 約 8 - 15 ms

這個結果很有用。

它告訴我們:下一步如果要變快,不是優先去改 MMA 指令。

因為 GEMM 那段已經不是最大塊。

真正該打的是:

DoubleDouble -> residue
residue result -> UInt384 CRT
host / device round trip
materialized residue matrix traffic

也就是說,下一步的 fused 版本應該長這樣:

load DoubleDouble tile
-> 在 GPU 上轉 scaled integer / residue
-> 對固定 modulus group 做 INT8 GEMM
-> accumulator 做 mod reduction
-> 做 partial CRT
-> 減少寫回 global memory 的中間結果

如果只是把 per-modulus GEMM 換成手寫 MMA,但 serialization 和 CRT 還留在 host,速度不會有本質變化。


這條路怎麼繼續往上

DoubleDouble 目前用的是 106-bit target。

如果未來想再往上,重點不是重寫概念,而是把這幾個東西加寬:

input precision bits
modulus count
CRT reconstruction width
output representation
oracle test

以 113-bit target、K=4096 來估:

required_crt_bits = 2 * 113 + 12 = 238 bits

目前 modulus table 前 32 個大約有:

240.4 bits

所以從 CRT 容量來看,往 113-bit target 走是可以的。

但今天先不把它叫完成。今天完成的是 DoubleDouble path,這是第一個站得住腳的高精度版本。


今天的結論

今天不是再講 CRT 原理,也不是再講 MMA 怎麼排。

今天真正做出來的是:

DoubleDouble input
-> scaled integer words
-> multi-modulus residue
-> INT8 GEMM per modulus
-> UInt384 CRT reconstruction
-> DoubleDouble output

這條路證明了:

只要 modulus count 和 reconstruction width 補夠,
FP64 看不到的低位可以被保留下來。

profile 也告訴我們下一步該做哪裡:

不是先改 MMA,
而是把 DoubleDouble serialization 和 CRT reconstruction 搬回 GPU,
再做真正的 fused residue / CRT pipeline。

速度還沒贏。

但精度往上推這件事,現在有具體版本了。

AI 還是在寫slob 但是今天我有多花時間迭代一下 同時也是我覺得至少我們可以做點什麼吧
雖然還是不能讓Nvidia emu滿意
但沒關西 生在有AI的時代真是太好了
可以保證沒工作 沒未來 只有一堆 債 要還
唉...(雖然其實跟AI也沒關係啦 是時機不好而已 沒事的)


上一篇
Day 12|速度比較與 Profile:我們到底差在哪 (重賽版)
下一篇
Day 14|把 DoubleDouble 再往上推:113-bit FP128 path (重賽版)
系列文
GPU 效能優化實戰:30 天從 Kernel 到 Profiling (重賽版)17
圖片
  熱門推薦
圖片
{{ item.channelVendor }} | {{ item.webinarstarted }} |
{{ formatDate(item.duration) }}
直播中

尚未有邦友留言

立即登入留言